Skip to content

[Diffusion] Rollout API: support return only SDE latents - #36994

Merged
Zhichenzzz merged 1 commit into
sgl-project:sglang-miles-h3from
Rockdu:rollout-traj-window
Aug 30, 2026
Merged

Zhichenzzz merged 1 commit into
sgl-project:sglang-miles-h3from
Rockdu:rollout-traj-window

Conversation

@Rockdu

@Rockdu Rockdu commented Aug 29, 2026 •

Copy link
Copy Markdown
Collaborator

Motivation

/rollout/generate returns the whole 0..T denoising trajectory, but flowGRPO training consumes only the SDE window: for each SDE step s it needs x_s, x_{s+1} and log_prob_s. On a 10-step schedule with a 2-step window that is 3 of the 11 latents; the other 8 can be thrown away, cutting trajectory payload size down by ~72%. These latents are fp32 and, for video models, dominate the response body — a MiniMax H3 t2va microgroup carries hundreds of MB of them.

rollout_return_step_indices already existed for this, but it was not usable:

  • the trainer could not tell which original step each surviving latent came from, so pairing x_s with x_{s+1} relied on the array being the full contiguous 0..T;
  • timesteps was filtered together with the latents, so the trainer lost the schedule entries it indexes by step number;
  • the MiniMax H3 path ignored the field entirely — its own collector recorded every step unconditionally.

Modifications

  • RolloutDitTrajectory gains latent_step_indices: the original step index of each returned latent (None = the full 0..T trajectory). Serialized alongside the trajectory and carried through the per-sample split.
  • RolloutDenoisingMixin: latents keep honoring rollout_return_step_indices and now record their provenance; timesteps is appended before the filter, so the schedule always comes back full-length. Scalar-sized arrays (timesteps, rollout_log_probs) are indexed by original step number, only the latents are filtered.
  • MiniMaxH3RolloutCollector: same filtering and provenance, so H3 stops ignoring the request; its timesteps now spans the full [T+1] including the terminal entry, matching sigmas and the generic path.

The consumer side lives in radixark/miles_diffusion#216: it requests S ∪ (S+1) for an SDE window S and pairs latents by the echoed provenance instead of by array position, so a non-contiguous window is correct by construction and a request/response disagreement raises instead of silently mispairing.

Accuracy Tests

Joint end-to-end runs with radixark/miles_diffusion#216 pinned to this branch — all e2e stages green (run 33244938347): stage-b-3-gpu-h200, stage-b-5-gpu-h200, stage-c-3-gpu-h200, stage-c-5-gpu-h200.

Every one of those recipes runs a step strategy, so all of them exercised the filtered transport, and each reproduced its pre-existing recorded metric standard bit for bit (strict comparison, --deterministic-mode):

Recipe SDE index set Result
test_sd3_ocr_grpo_2xGPU contiguous window bitwise identical to standard
test_qwenimage_pickscore_grpo_5xGPU contiguous window bitwise identical to standard
test_ltx23_pickscore_grpo_4xGPU non-contiguous candidates bitwise identical to standard
test_wan22_pickscore_grpo_17xGPU_single_node_4xGPU_proxy non-contiguous candidates bitwise identical to standard

Two further e2e CIs are being added for the models this PR touches. Each recorded its standard off this branch, then reproduced it on top of it:

New e2e SDE index set Result on this branch
test_h3_t2va_grpo_2xGPU (miles_diffusion#217) non-contiguous candidates standard recorded with the stock sglang-miles-h3 engine; all 8 metric series identical to the last digit, and that baseline passes in CI (run 33271965673, 1490 s)
test_cosmos3_pickscore_grpo_t2i_4xGPU (miles_diffusion#219) window 8-11, does not start at step 0 standard recorded on sglang main; all 10 metric series bitwise identical here (run). On the stock engine the same run instead fails in rollout to train conversion: trajectory lacks latents for steps [10, 11, 12] (have [0, 1, 2]) (run)

cosmos3 is the case that cannot pass without the provenance echo: its Karras grid puts the useful window at steps 8-11, so array position never coincides with step number and pairing by position would silently pair x_10 with x_0. Its train/model_output_{mean_abs_diff,rel_max} compare the raw DiT outputs between engine and trainer, so the equality covers the tensors themselves rather than a reward summary.

Checklist

  • Format your code according to the Format code with pre-commit. — run on the touched files; black reformatting is included in the commit
  • Add unit tests according to the Run and add unit tests. — no new unit test; the filter and provenance were verified behaviourally against the H3 collector and the generic mixin, and by the bitwise e2e comparisons above
  • Update documentation according to Write documentations. — not needed, the field is described where it is declared
  • Provide accuracy and speed benchmark results according to Test the accuracy and Benchmark the speed.
  • Follow the SGLang code style guidance.

CI States

Latest PR Test (Base): ❌ Run #33282596522
Latest PR Test (Extra): ❌ Run #33282596363
Latest PR Test (AMD ROCm 7.2): ❌ Run #33282596623

@github-actions github-actions Bot added the diffusion SGLang Diffusion label Aug 29, 2026
@Rockdu Rockdu changed the title [Diffusion] Rollout API: filter trajectory latents to the requested w… [Diffusion] Rollout API: filter trajectory latents to the requested step-index Aug 29, 2026
@Rockdu
Rockdu force-pushed the rollout-traj-window branch from c50aebb to e9bea69 Compare August 29, 2026 09:02
@Rockdu
Rockdu force-pushed the rollout-traj-window branch from e9bea69 to cbeeeb5 Compare August 30, 2026 00:08
@Rockdu Rockdu changed the title [Diffusion] Rollout API: filter trajectory latents to the requested step-index [Diffusion] Rollout API: filter only SDE latents Aug 30, 2026
@Rockdu Rockdu changed the title [Diffusion] Rollout API: filter only SDE latents [Diffusion] Rollout API: support filter only SDE latents Aug 30, 2026
@Rockdu Rockdu changed the title [Diffusion] Rollout API: support filter only SDE latents [Diffusion] Rollout API: support return only SDE latents Aug 30, 2026
@Rockdu
Rockdu marked this pull request as ready for review August 30, 2026 00:50
@Zhichenzzz
Zhichenzzz merged commit c80d803 into sgl-project:sglang-miles-h3 Aug 30, 2026
81 of 90 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

diffusion SGLang Diffusion

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants